Skip to content

ZeRO 显存优化系列 ​

标签
AI/infra/显存管理
AI/infra/并行策略
字数
1496 字
阅读时间
7 分钟

02-分布式训练总论与显存账本 算过一笔账:BF16 + Adam 下每参数静态显存约 16Ψ,7B 模型就要 112 GB。ZeRO(Zero Redundancy Optimizer)是 DeepSpeed 的核心技术,它从「冗余在哪里」这个问题出发,逐阶段切掉这些冗余。

冗余分析 ​

DDP 里每张卡都冗余存储完整的三份东西:

组成每卡存储
优化器状态12Ψ
梯度2Ψ
参数2Ψ

N 卡集群里,N 份完全相同的副本 —— 其中 (N−1)/N 的存储是纯粹浪费的。

设计思想是「切分-聚合」范式:平时每卡只存 1/N,需要时通过通信获取完整数据,用完即弃。

三阶段的顺序从最容易切的开始(优化器状态),到最难切的(参数)—— 因为越往前切,通信代价越小。

三个阶段 ​

ZeRO-1:只切优化器状态 ​

切什么:Adam 的 FP32 参数副本 + 一阶动量 + 二阶动量,共 12Ψ。

通信:梯度仍需 AllReduce(与 DDP 相同),参数更新后需要 AllGather 同步参数。

每卡显存:16Ψ⟶4Ψ+12ΨN

这一刀最划算的原因:Adam 状态占了 16Ψ 里的 12(见 02-分布式训练总论与显存账本),切它等于切掉了大头,而通信只增加一次参数 AllGather。

ZeRO-2:再切梯度 ​

切什么:在 ZeRO-1 基础上,梯度也按分片存储。

通信模式的关键变化:反向传播时用 ReduceScatter 替代 AllReduce —— 每卡只保留自己负责分片的聚合梯度。

每卡显存:16Ψ⟶2Ψ+14ΨN通信量:ReduceScatter=Ψ(AllReduce 的 2Ψ 的一半)

注意这里有个容易漏的点:ReduceScatter 的 Ψ 比 AllReduce 的 2Ψ「少一半」,但 AllReduce 隐含的那次 AllGather 并没有消失 —— 只是被推迟到了参数更新时。通信量的对比要按完整一步算,不能只比单次调用。

ZeRO-3:连参数也切 ​

切什么:参数、梯度、优化器状态全部切分。

通信模式:

前向:AllGather 重组当前层参数 → 计算 → 释放
反向:AllGather 参数 → 计算梯度 → ReduceScatter 梯度 → 释放参数
每卡显存:16Ψ⟶16ΨN(理想线性缩放)通信量:3Ψ(Ψ⏟前向 AllGather+Ψ⏟反向 AllGather+Ψ⏟ReduceScatter),比DDP多50%

ZeRO-3 与前两个阶段有性质上的差别:ZeRO-1/2 是「省显存,通信代价小」,ZeRO-3 是用 50% 的额外通信换线性显存缩放。它只在「参数本身单卡装不下」时才值得 —— 这是选型的分水岭。

每张卡上到底存了什么(N 卡,Ψ 是参数量):

  DDP:三份都全量冗余

    ┌──────────────┬──────────┬──────────────────────────────────┐
    │ 参数 2Ψ      │ 梯度 2Ψ  │ 优化器状态 12Ψ(含 FP32 主权重)   │  ⇒ 16Ψ/卡
    └──────────────┴──────────┴──────────────────────────────────┘

  ZeRO-1:只切优化器状态

    ┌──────────────┬──────────┬──────────────────────────────────┐
    │ 参数 2Ψ      │ 梯度 2Ψ  │ 优化器 12Ψ ÷ N                    │  ⇒ 4Ψ + 12Ψ/N
    └──────────────┴──────────┴──────────────────────────────────┘

  ZeRO-2:再切梯度

    ┌──────────────┬──────────┬──────────────────────────────────┐
    │ 参数 2Ψ      │ 梯度 2Ψ÷N│ 优化器 12Ψ ÷ N                    │  ⇒ 2Ψ + 14Ψ/N
    └──────────────┴──────────┴──────────────────────────────────┘

  ZeRO-3:参数也切

    ┌──────────────┬──────────┬──────────────────────────────────┐
    │ 参数 2Ψ ÷ N  │ 梯度 2Ψ÷N│ 优化器 12Ψ ÷ N                    │  ⇒ 16Ψ/N
    └──────────────┴──────────┴──────────────────────────────────┘

  ⇒ 每往下一个阶段多切一样东西,通信代价随之递增;
     ZeRO-3 相比 DDP 的通信量多 50%,换来的是线性显存缩放。

Offload:把状态搬到 CPU 和 SSD ​

方案做法适用代价
ZeRO-Offload优化器状态与梯度计算卸载到 CPU,GPU 只做前向/反向少卡(1–4 卡)训大模型PCIe 带宽成瓶颈(PCIe 4.0 约 32 GB/s)
ZeRO-Infinity在 Offload 基础上进一步利用 NVMe SSD万亿参数在有限 GPU 上训练I/O 带宽;靠分块预取(prefetch)与计算-I/O 重叠掩盖

Offload 的本质是「拿一条慢得多的链路换容量」:PCIe 4.0 的 32 GB/s 对比 NVLink 4.0 的 900 GB/s(见 03-多卡互联与集群网络)—— 差 28 倍。所以它只适合「少卡、无别的选择」的场景,不是通用优化。

选型 ​

阶段省什么通信代价什么时候用
ZeRO-1优化器状态(12ΨN)与 DDP 接近参数 + 梯度单卡装得下
ZeRO-2+ 梯度略高于 DDP参数单卡装得下
ZeRO-3全部(16ΨN)比 DDP 多 50%参数也装不下
Offload用 CPU 内存+ PCIe 传输少卡大模型

判据是「装不下的到底是哪一部分」 —— 优化器状态装不下用 1,梯度也装不下用 2,参数都装不下才用 3。

判据是「装不下的到底是哪一部分」,与模型多大无关:

   优化器状态装不下 ──▶ ZeRO-1     通信与 DDP 接近
        │
   梯度也装不下     ──▶ ZeRO-2     通信略高于 DDP
        │
   参数都装不下     ──▶ ZeRO-3     通信比 DDP 多 50%
        │
   CPU 内存也要用上 ──▶ Offload    拿 PCIe 的 32 GB/s 换容量
                                   (对比 NVLink 4.0 的 900 GB/s,差 28 倍)

ZeRO 与 FSDP 的对应 ​

ZeROPyTorch FSDP
ZeRO-2SHARD_GRAD_OP
ZeRO-3FULL_SHARD

三者的分工:

  • ZeRO 是算法层面的概念(Microsoft, arXiv:1910.02054)
  • FSDP 是 PyTorch 原生实现
  • DeepSpeed 是微软的独立实现

怎么选:PyTorch 生态内优先 FSDP(原生、与 torch.compile / torch.distributed 集成好);需要 Offload / Infinity,或已有 DeepSpeed 配置时用 DeepSpeed。

相关 ​

参考 ​

贡献者 ​

文件历史 ​